代码实现 RMSNorm
从零实现 RMSNorm:一行公式如何变成 PyTorch 代码
先回忆下 RMSNorm(Root Mean Square Normalization)的原理:
以及 RMSNorm 的计算公式:
其中:
:输入向量 :均方根 :可学习参数(weight)
接下来将结合一段真实的大模型代码,逐行讲解 RMSNorm 的实现原理。
一、完整代码
class RMSNorm(torch.nn.Module):
def __init__(self, dim: int, eps: float = 1e-5):
super().__init__()
self.eps = eps
self.weight = torch.nn.Parameter(torch.ones(dim))
def norm(self, x):
return x * torch.rsqrt(
x.pow(2).mean(-1, keepdim=True) + self.eps
)
def forward(self, x):
return (self.weight * self.norm(x.float())).type_as(x)
这段代码非常短。
真正的核心逻辑只有一行:
x * torch.rsqrt(
x.pow(2).mean(-1, keepdim=True) + self.eps
)
理解这一行,基本就理解了 RMSNorm。
二、继承 torch.nn.Module
class RMSNorm(torch.nn.Module):
PyTorch 中所有神经网络层都继承自:
torch.nn.Module
例如:
nn.Linear
nn.Conv2d
nn.LayerNorm
本质上也都是 Module。
这样做的好处是:
- 自动管理参数
- 自动计算梯度
- 支持模型保存和加载
- 支持 GPU 加速
例如:
model = RMSNorm(512)
for name, param in model.named_parameters():
print(name)
输出:
weight
说明 RMSNorm 中存在一个可学习参数。
三、初始化函数
init()
def __init__(self, dim: int, eps: float = 1e-5):
参数:
| 参数 | 含义 |
|---|---|
| dim | 隐藏层维度 |
| eps | 防止除零的小常数 |
例如:
RMSNorm(4096)
表示:
hidden_size = 4096
这是 LLaMA、Qwen 等模型中常见的隐藏维度。
四、为什么要有 eps?
self.eps = eps
默认:
eps = 1e-5
原因很简单。
假设输入:
x = [0,0,0]
那么:
此时会出现:
程序直接崩溃。
因此需要:
RMS(x) + eps
变成:
避免除零错误。
五、可学习参数 weight
torch.ones
torch.ones(dim)
例如:
torch.ones(4)
输出:
tensor([1., 1., 1., 1.])
表示创建一个全是 1 的向量。
torch.nn.Parameter
self.weight = torch.nn.Parameter(
torch.ones(dim)
)
等价于:
γ = [1,1,1,...]
这正对应 RMSNorm 公式中的:
为什么需要 weight?
如果没有 weight:
模型表达能力会下降。
因此加入:
让模型自己学习:
- 哪些维度放大
- 哪些维度缩小
这就是:
self.weight
的作用。
六、进入核心:norm()
def norm(self, x):
这里真正实现 RMSNorm。
第一步:平方
x.pow(2)
例如:
x = [2,4,6]
得到:
[4,16,36]
对应数学公式:
第二步:求均值
.mean(-1, keepdim=True)
继续上面的例子:
[4,16,36]
均值:
得到:
18.67
对应公式:
为什么是 -1?
mean(-1)
表示:
最后一个维度
例如:
x.shape
=
(2,3)
[
[1,2,3],
[4,5,6]
]
执行:
x.mean(-1)
得到:
[
2,
5
]
即:
每个 token 单独计算 RMS。
keepdim=True
如果:
x.shape=(2,3)
那么:
x.mean(-1)
结果:
shape=(2,)
变成:
[2,5]
而:
x.mean(-1, keepdim=True)
结果:
shape=(2,1)
即:
[
[2],
[5]
]
这样方便后续广播运算(Broadcast)。
七、加上 eps
x.pow(2).mean(-1, keepdim=True)
+ self.eps
对应:
避免分母为 0。
八、torch.rsqrt()
这是 RMSNorm 中最容易被忽略的一步。
rsqrt 是什么?
torch.rsqrt(x)
表示:
例如:
torch.rsqrt(
torch.tensor([1.,4.,9.])
)
输出:
[1.0,0.5,0.3333]
因为:
为什么不用 sqrt?
RMSNorm 需要:
而:
torch.rsqrt(...)
直接得到:
于是:
x * rsqrt(...)
就等价于:
少做一次除法。
GPU 上更高效。
九、完成 RMS 归一化
这一句:
x * torch.rsqrt(
x.pow(2).mean(-1, keepdim=True)
+ self.eps
)
完整对应:
这就是 RMSNorm 的核心数学公式。
十、forward()
def forward(self, x):
PyTorch 约定:
output = model(x)
实际上调用:
model.forward(x)
x.float()
self.norm(x.float())
很多大模型使用:
float16
bfloat16
训练。
例如:
dtype=torch.float16
为了提高数值稳定性:
归一化时先转成:
float32
即:
x.float()
这样:
RMS
计算更加准确。
乘以 weight
self.weight * self.norm(...)
对应数学公式:
其中:
type_as(x)
.type_as(x)
表示:
结果转回原来的数据类型
例如:
输入:
float16
归一化时:
float32
最终再转回:
float16
这样:
- 保持数值稳定
- 节省显存
这是现代 LLM 的标准做法。
十一、完整执行流程
假设输入:
x = [2,4,6]
Step1:平方
[4,16,36]
Step2:求均值
18.67
Step3:加 eps
18.67001
Step4:rsqrt
1 / sqrt(18.67001) ≈ 0.231
Step5:乘回原向量
[
0.462,
0.924,
1.386
]
Step6:乘 weight
假设:
weight = [1,1,1]
结果不变:
[
0.462,
0.924,
1.386
]
如果训练后:
weight =
[
1.2,
0.8,
1.5
]
结果变成:
[
0.554,
0.739,
2.079
]
模型便可以自主调整每个维度的重要性。
十二、一句话总结
RMSNorm 的实现本质上只有一句核心代码:
x * torch.rsqrt(
x.pow(2).mean(-1, keepdim=True) + eps
)
它对应数学公式:
然后再乘上可学习参数:
相比 LayerNorm,RMSNorm 去掉了均值计算和方差计算,只保留尺度归一化(Scaling),因此实现更简单、计算更高效,也成为 LLaMA、Qwen、DeepSeek 等现代大模型的默认归一化方案。